AdamWeightDecay
对权重执行 Adam Weight Decay(AdamW)优化更新,原地写回 var、m、v。
\[\begin{split}\begin{aligned}
m_t &= m_{t-1} + (g_t - m_{t-1}) \cdot (1 - \beta_1) \\
v_t &= v_{t-1} + (g_t^{2} - v_{t-1}) \cdot (1 - \beta_2) \\
var_t &= var_{t-1} - lr \cdot \left( \frac{m_t}{\sqrt{v_t} + \epsilon} + decay \cdot var_{t-1} \right)
\end{aligned}\end{split}\]
其中 \(g_t\) 为梯度,\(m_t\) / \(v_t\) 为一阶 / 二阶动量,
\(decay\) 为权重衰减系数。本实现对区间 [start, end) 内元素逐点更新。
- 输入:
var - 待更新权重张量地址(输入/输出,原地更新)
m - 一阶动量张量地址(输入/输出,原地更新)
v - 二阶动量张量地址(输入/输出,原地更新)
grad - 梯度张量地址
float_Parameters - 浮点超参数组,长度 5,布局见下
int_Parameters - 整数参数数组,长度 2,布局见下
core_mask - 核掩码(仅共享存储版本使用)
float_Parameters 布局:
[0]lr- 学习率
[1]beta1- 一阶动量衰减系数
[2]beta2- 二阶动量衰减系数
[3]epsilon- 数值稳定项
[4]decay- 权重衰减系数
int_Parameters 布局:
[0]start- 参与计算的起始索引(含)
[1]end- 参与计算的结束索引(不含)
- 输出:
var / m / v - 原地写回更新结果
- 支持平台:
FT78NEMT7004
备注
FT78NE 支持 fp32
MT7004 支持 fp16、fp32
hp_版本中float_Parameters元素类型为float16;fp_版本为float
共享存储版本:
-
void hp_adam_weight_decay_s(float16 *var, float16 *m, float16 *v, float16 *grad, float16 *float_Parameters, int *int_Parameters, int core_mask)
-
void fp_adam_weight_decay_s(float *var, float *m, float *v, float *grad, float *float_Parameters, int *int_Parameters, int core_mask)
C调用示例:
1// MT7004 示例(共享存储多核,DDR 地址)
2void TestAdamWeightDecaySMCFp32(int length, int core_mask) {
3 int core_id = get_core_id();
4 int logic_core_id = GetLogicCoreId(core_mask, core_id);
5 int core_num = GetCoreNum(core_mask);
6 float *var = (float *)0x81000000;
7 float *m = (float *)0x82000000;
8 float *v = (float *)0x83000000;
9 float *grad = (float *)0x84000000;
10 float *float_Parameters = (float *)0x85000000;
11 int *int_Parameters = (int *)0x85000014;
12 if (logic_core_id == 0) {
13 float_Parameters[0] = 0.001f; // lr
14 float_Parameters[1] = 0.9f; // beta1
15 float_Parameters[2] = 0.999f; // beta2
16 float_Parameters[3] = 1e-4f; // epsilon
17 float_Parameters[4] = 0.01f; // decay
18 int_Parameters[0] = 0; // start
19 int_Parameters[1] = length; // end
20 }
21 sys_bar(0, core_num);
22 fp_adam_weight_decay_s(var, m, v, grad, float_Parameters, int_Parameters, core_mask);
23}
24
25void main() {
26 int core_mask = 0b1111;
27 TestAdamWeightDecaySMCFp32(4096, core_mask);
28}
私有存储版本:
-
void hp_adam_weight_decay_p(float16 *var, float16 *m, float16 *v, float16 *grad, float16 *float_Parameters, int *int_Parameters)
-
void fp_adam_weight_decay_p(float *var, float *m, float *v, float *grad, float *float_Parameters, int *int_Parameters)
C调用示例:
1// MT7004 示例(私有存储单核,AM 地址)
2void TestAdamWeightDecayAMFp32(int length) {
3 float *var = (float *)0x10000000;
4 float *m = (float *)0x10020000;
5 float *v = (float *)0x10040000;
6 float *grad = (float *)0x10060000;
7 float *float_Parameters = (float *)0x1007F000;
8 int *int_Parameters = (int *)0x1007F014;
9 float_Parameters[0] = 0.001f; // lr
10 float_Parameters[1] = 0.9f; // beta1
11 float_Parameters[2] = 0.999f; // beta2
12 float_Parameters[3] = 1e-4f; // epsilon
13 float_Parameters[4] = 0.01f; // decay
14 int_Parameters[0] = 0; // start
15 int_Parameters[1] = length; // end
16 fp_adam_weight_decay_p(var, m, v, grad, float_Parameters, int_Parameters);
17}
18
19void main() {
20 TestAdamWeightDecayAMFp32(2048);
21}